昨天我們把題目換掉了(換題目水 嘿嘿)。
一樣 先ai slob 然後我再看時間修 。 總之 先勉強看個
這次不是 GEMM,而是:
y = H^{\otimes N} x
如果把 H^ N 展開成矩陣,那會是一個非常大的 dense matrix;但真正要算的不是 dense GEMM,而是一層一層的 butterfly。
今天先看第一段切分:當 N 還不大時,能不能把整個 FWT chunk 放在一個 CUDA block 裡面做完?
答案是可以,但 kernel 不能寫成「每個 stage 都讀寫 global memory」。那樣會把 FWT 最糟糕的一面放大:每一層只做一次加減,卻一直搬資料。
所以這個版本的方向是:
global memory 只進出一次
中間的 butterfly 盡量留在 register / warp / shared memory
FWT 的 vector 長度是 2^N。
假設用 FP32,一個 element 是 4 bytes:
N=10: 2^10 floats = 4 KB
N=12: 2^12 floats = 16 KB
N=14: 2^14 floats = 64 KB
N=15: 2^15 floats = 128 KB
這就是 single-block kernel 的第一個邊界。
N=14 時,整個 chunk 是 64 KB。這已經很接近 shared memory 壓力,但還能靠 register blocking 和 shared memory exchange 撐住。
N=15 時,整個 chunk 是 128 KB。如果你想把完整資料丟進 shared memory,會很不舒服,甚至直接撞限制。
所以真正的寫法不是:
把 128 KB state 全部放進 shared memory
而是:
每個 thread 拿一部分資料放 register
shared memory 只拿來做跨 warp 的資料重排
這是 Day 16 的核心。
case_3 裡的 native path 主要在:
src/ImplicitHadamardNative.cu
single-block 的核心 kernel 是:
template<int N, int THREADS, int NUM_B>
__global__ void native_fp32_extreme_fwt_kernel(
const float* __restrict__ in_state,
float* __restrict__ out_state,
float norm_factor
)
這裡的 NUM_B 很重要。
每個 thread 不只處理一個 float,而是處理多個 float4:
float4 reg[NUM_B];
也就是:
每個 thread 持有 NUM_B 組 float4
每組 float4 = 4 個 FP32
所以一個 block 處理的元素數量是:
THREADS * NUM_B * 4
例如 N=14:
THREADS = 256
NUM_B = 16
256 * 16 * 4 = 16384 = 2^14
這剛好是一整個 N=14 chunk。
每個 float4 裡面有 4 個值:
v.x, v.y, v.z, v.w
FWT 的前兩層 stride 很小,所以不需要跨 thread。
可以直接在同一個 thread 裡做:
float a0 = reg[b].x + reg[b].y;
float b0 = reg[b].x - reg[b].y;
float a1 = reg[b].z + reg[b].w;
float b1 = reg[b].z - reg[b].w;
reg[b].x = a0 + a1;
reg[b].y = b0 + b1;
reg[b].z = a0 - a1;
reg[b].w = b0 - b1;
這裡沒有 shared memory,也沒有 warp shuffle。
資料在 register 裡,ALU 直接做加減。
這也是為什麼一開始用 float4 很自然:它不是為了好看,而是剛好可以把前兩層 butterfly 吃掉。
接著 stride 會跨 thread,但還在同一個 warp 裡。
FWT 的配對規則是:
peer = lane xor stride
所以 CUDA 裡可以直接用:
float peer = __shfl_xor_sync(0xffffffff, value, stride);
對 float4 來說,就是四個 lane value 都各自 shuffle:
peer.x = __shfl_xor_sync(0xffffffff, reg[b].x, stride_thread);
peer.y = __shfl_xor_sync(0xffffffff, reg[b].y, stride_thread);
peer.z = __shfl_xor_sync(0xffffffff, reg[b].z, stride_thread);
peer.w = __shfl_xor_sync(0xffffffff, reg[b].w, stride_thread);
然後根據目前 lane 在 pair 的上半部或下半部,決定是:
上半部:self + peer
下半部:peer - self
這段最關鍵的地方是:warp shuffle 是 register-to-register 的資料交換。
它沒有把資料寫回 global memory。
對 FWT 這種每次只做 add/sub 的 workload,這比多幾個 FLOPS 更重要。
到了更大的 stride,資料會跨 warp。
warp shuffle 只能處理同一個 warp 的 lane exchange,所以這時候需要 shared memory。
但這裡要小心:shared memory 不應該變成主要計算場所。
比較好的用法是:
register -> shared memory 做 layout exchange
shared memory -> register
接著繼續用 warp shuffle / register butterfly
這就是後來 N=15 版本最重要的洞見。
原本直覺可能會覺得:
N=15 有 128 KB,所以需要 128 KB shared memory
但實際上可以改成:
128 KB 的資料分散在每個 thread 的 register
shared memory 只拿來做 corner turning
這個差別非常大。
case_3 裡後來有一個更激進的版本:
native_fp32_extreme_fwt_kernel_v2<15, 256, 32>
這個配置代表:
THREADS = 256
NUM_B = 32
256 * 32 * 4 = 32768 = 2^15
也就是一個 block 直接吃完整個 N=15 chunk。
每個 thread 持有:
32 個 float4 = 128 個 float
這會帶來很高的 register 壓力,但它換來一件事:
N=15 不需要切成兩個 kernel
中間不需要完整寫回 global memory 再讀回來
這正是 N=15 能不能快的關鍵。
當資料跨 warp 交換時,如果大家用很規則的方式寫 shared memory,很容易打到同一批 bank。
shared memory bank conflict 會讓看起來很近的資料交換變慢。
所以 N=15 版本用了 XOR swizzle 的概念:
smem index 裡混入 lane_id xor warp_id
概念上就是讓不同 warp/lane 的資料不要全部撞到同一個 bank。
文章先不把 swizzle 寫成完整推導;這部分後面如果要單獨拆,可以再把 bank mapping 畫出來。
今天只需要掌握一件事:
shared memory 在這裡不是 cache,
而是一個可控的 cross-warp exchange buffer。
native path 的 dispatch 大致像這樣:
N=7 -> native_fp32_extreme_fwt_kernel<7, 32, 1>
N=8 -> native_fp32_extreme_fwt_kernel<8, 64, 1>
N=9 -> native_fp32_extreme_fwt_kernel<9, 128, 1>
N=10 -> native_fp32_extreme_fwt_kernel<10,256, 1>
N=11 -> native_fp32_extreme_fwt_kernel<11,256, 2>
N=12 -> native_fp32_extreme_fwt_kernel<12,256, 4>
N=13 -> native_fp32_extreme_fwt_kernel<13,256, 8>
N=14 -> native_fp32_extreme_fwt_kernel<14,256, 16>
N=15 -> native_fp32_extreme_fwt_kernel_v2<15,256,32>
這個表其實就是一個容量表。
N 每增加 1,資料量翻倍。當 thread 數固定在 256 時,就只能讓每個 thread 拿更多 float4。
所以 NUM_B 會一路增加:
1, 2, 4, 8, 16, 32
這不是 tuning magic,而是 2^N 的結果。
這種 kernel 最容易出現的錯不是 crash,而是 silent wrong answer。
尤其是 N=15 的 exchange,如果 swizzle 或讀回順序錯了,程式照樣跑完,但 butterfly pair 已經亂掉。
所以驗證不能只看 norm 差不多。
至少要做:
1. 小 N 對 naive PyTorch FWT
2. N=15 對 Dao-AILab fast-hadamard-transform
3. 檢查 normalization 是否一致
4. 看 max diff,而不是只看平均誤差
case_3 的結果裡,N=15 有一個很重要的數字:
QDP Native FP32: 82.98 us
Dao FHT: 98.35 us
max diff: 0.0e+00
0.0e+00 不是說 floating point 永遠沒有誤差,而是代表在這個 benchmark 設定下,計算順序跟比較標的對齊了。
這對 FWT 很重要,因為不同 butterfly 順序會讓 FP32 rounding path 不同。
single-block FWT 的重點不是把算力打滿。
重點是讓資料不要離開晶片:
thread-local:float4 內部 butterfly
warp-local:__shfl_xor_sync
cross-warp:shared memory swizzle
large stage:register blocking
N=14 是比較自然的 shared-memory/register 邊界。
N=15 則需要更激進:每個 thread 拿 32 個 float4,用 register space 承擔完整 chunk,再用 shared memory 做 corner turning。
下一篇要看的是:如果 N > 15,一個 block 放不下了,這時候要怎麼切。